Vamana build optimization set and fp16 support - #2264
Conversation
(cherry picked from commit 14e36b3)
(cherry picked from commit d8b547b)
(cherry picked from commit 4d970e0)
…L2 comparators. (cherry picked from commit 0746d15)
…e vector in shared memory in the RobustPrune occlusion loop (cherry picked from commit f45fd1b49283434eb4a3017da069ead501e938c3)
…efix sum with cub scan and hoist per-batch reverse-edge allocations (cherry picked from commit 3b8650f5ebea52421f187332a2f6f3bdd599c42e)
…lusion across multiple warps per query (raise occupancy) (cherry picked from commit 2e02f938f97e65ca073daf07397d538432b52867)
… query->existing-edge distances in the RobustPrune merge (avoid recompute) (cherry picked from commit 041d355c585b7f98302d054b80ac287d637a07e4)
…oords in FP16 smem for dim>=512 to raise GreedySearch occupancy (salvage of N8)
…nce (one warp) instead of redundantly on all 128 threads, then broadcast
…ck (4 vs 8) to raise occupancy/MLP on the degree-64 occlusion sweep
| #define KERNEL_TIMING (RAFT_LOG_ACTIVE_LEVEL <= RAPIDS_LOGGER_LOG_LEVEL_DEBUG) | ||
|
|
||
| template <typename accT, typename IdxT> | ||
| __global__ void gather_query_sizes(QueryCandidates<IdxT, accT>* query_list, |
There was a problem hiding this comment.
Do you have a sense for how much this is adding to the binary size? @divyegala should be able to explain how to see the deployed metrics for PRs.
|
/ok to test c055c60 |
|
/ok to test 350ed89 |
@jamxia155, there was an error processing your request: See the following link for more information: https://docs.gha-runners.nvidia.com/cpr/e/2/ |
|
/ok to test 74caba9 |
|
/ok to test b8238f3 |
|
|
||
| num_neighbors = degree; | ||
| __syncthreads(); | ||
| num_neighbors[warpIdx] = degree; |
There was a problem hiding this comment.
maybe make only lane 0 write to this?
There was a problem hiding this comment.
Having just lane 0 write it would require a sync after. Seems simpler/cleaner to just have all threads write, since they have the same value.
| __float2half(0.0f), __float2half(0.0f), __float2half(0.0f), __float2half(0.0f)}; | ||
| for (int i = threadIdx.x; i < src_vec->Dim; i += 4 * blockDim.x) { | ||
| temp_dst[0] = dst_vec->coords[i]; | ||
| if (i + 32 < src_vec->Dim) temp_dst[1] = dst_vec->coords[i + 32]; |
There was a problem hiding this comment.
Should we use blockDim.x instead of hardcoding to 32 here?
There was a problem hiding this comment.
I agree with this. We should not be hardcoding things because the kernel invocation can change (or the same kernel may be called somewhere else in the future).
There was a problem hiding this comment.
This is run by a single warp within the block, and uses a warp width stride. I can change it to a define macro if that helps?
There was a problem hiding this comment.
Replaced with raft::WarpSize
tarang-jain
left a comment
There was a problem hiding this comment.
There are a lot of changes in this PR where I want to basically tell the exact same things -- using raft macros and raft primtives wherever possible. We should certainly avoid writing small kernels which can easily be a lambda for raft::linalg. It helps keeps things readable and binary size increases can become slightly more predictable. For example, earlier we have found that raft::linalg::map kernels generally add less to the binary size than thrust::for_each (or similar thrust primitives).
| Node<SUMTYPE> input_data, | ||
| SUMTYPE* cur_max_val, | ||
| int* max_idx) | ||
| __inline__ __device__ void parallel_pq_max_enqueue_warp(Node<SUMTYPE>* pq, |
There was a problem hiding this comment.
Nit: we have macros to mark kernels: such as _RAFT_INLINE
There was a problem hiding this comment.
I wouldnt say these macros should block the PR from merging.
There was a problem hiding this comment.
I can change this throughout, but I can't find the docs in raft on the different INLINE macros.
| __global__ void scatter_prefix_offsets(QueryCandidates<IdxT, accT>* query_list, | ||
| const int* edge_offsets, | ||
| int count) | ||
| { |
There was a problem hiding this comment.
same here: we should be able to achieve these with raft primitives.
There was a problem hiding this comment.
This scatter kernel is not so simple. I could do it but would need to re-work the code a bit more to use a strided mdspan or something. Can work on this if it's a blocker for the PR.
There was a problem hiding this comment.
We can do this as a follow-up. @bkarsin do you mind creating an issue and referencing it in a comment here just to get this merged?
There was a problem hiding this comment.
Created issue #2493 detailing that this kernel should be replaced with RAFT primitives.
tarang-jain
left a comment
There was a problem hiding this comment.
I have not looked into the implementation itself (what the code inside the kernels does) and my review continues to be along the same direction as my previous comments -- to improve overall code quality and readability. Pointing out a couple of very small things.
|
|
||
| const int laneId = threadIdx.x & 31; | ||
| const int warpId = threadIdx.x >> 5; | ||
| const int num_warps = blockDim.x >> 5; |
There was a problem hiding this comment.
Lets keep this simple. I have been advised by @tfeher in the past to not use bit arithmetic for simple division / mod operations, since the compiler is typically smart enough to automatically simplify these. So we can do again things like laneId = threadIdx.x % raft::WarpSize and warpId = threadIdx.x / 32.
There was a problem hiding this comment.
I think this can also be ignored @bkarsin . Its too much nit-picking.
| auto scan_temp_storage = raft::make_device_mdarray<uint8_t>( | ||
| res, large_ws, raft::make_extents<int64_t>(std::max(scan_temp_bytes, size_t{1}))); | ||
|
|
||
| thrust::device_vector<IdxT> edge_dest_vec(max_total_edges); |
There was a problem hiding this comment.
Use raft::device_vector here
|
/ok to test e1e0b7e |
|
/ok to test f0579f9 |
| The following distance metrics are supported: | ||
| - L2Expanded | ||
| - L2Expanded (sqeuclidean) | ||
| - L2SqrtExpanded (l2 / euclidean distance) |
There was a problem hiding this comment.
Also update the C doc for cuvsVamanaBuild for the L2SqrtExpanded distance
| __device__ SUMTYPE l2_ILP2_half(Point<__half, SUMTYPE>* src_vec, Point<__half, SUMTYPE>* dst_vec) | ||
| { | ||
| __half temp_dst[2] = {__float2half(0.0f), __float2half(0.0f)}; | ||
| __half partial_sum[2] = {__float2half(0.0f), __float2half(0.0f)}; | ||
| for (int i = threadIdx.x; i < src_vec->Dim; i += 2 * blockDim.x) { | ||
| temp_dst[0] = dst_vec->coords[i]; | ||
| if (i + raft::WarpSize < src_vec->Dim) temp_dst[1] = dst_vec->coords[i + raft::WarpSize]; | ||
|
|
||
| l2_half_fma_sq(partial_sum[0], src_vec[0].coords[i], temp_dst[0]); | ||
| if (i + raft::WarpSize < src_vec->Dim) | ||
| l2_half_fma_sq(partial_sum[1], src_vec[0].coords[i + raft::WarpSize], temp_dst[1]); | ||
| } | ||
| partial_sum[0] = __hadd(partial_sum[0], partial_sum[1]); | ||
|
|
||
| return l2_half_warp_reduce<SUMTYPE>(partial_sum[0]); | ||
| } |
There was a problem hiding this comment.
Would it make sense to also provide an overload of l2_half_fma_sq with half2 type? I see that __hfma2 and __hsub2 exist.
There was a problem hiding this comment.
It might benefit l2_ILP4_half below too
There was a problem hiding this comment.
Good point, though this is not a super simple change. Would have to choose what versions to replace with this (only ILP variants?) and deal with edge cases like odd dimensions. Would also require some re-benchmarking. To avoid making this PR even bigger, would adding at issue for this be better?
There was a problem hiding this comment.
Okay then can you create that issue, and add a comment in the code linking to that issue for better tracking please?
| __device__ __forceinline__ void l2_fma_sq2(SUMTYPE& acc, float sx, float sy, float2 dst2) | ||
| { | ||
| float dx = sx - dst2.x; | ||
| float dy = sy - dst2.y; | ||
| acc = fmaf(dx, dx, acc); | ||
| acc = fmaf(dy, dy, acc); | ||
| } |
There was a problem hiding this comment.
This might also benefit from fsub2/fmaf2. The data loaded is already float2
Series of GPU Vamana build performance optimizations that addresses #2178 and #1757. Initial estimates from #1757 were not accurate, so many other optimizations were tried (some abandoned, some successful). This PR includes:
GreedySearch optimizations:
RobustPrune optimizations:
General optimizations:
Together these optimizations give significant speedups across all configs with minimal recall variance compared to the current baseline. I benchmarked performance across a range of synthetic datasets and two real-world datasets. (NOTE: finishing benchmarks and will update tables below once they are all collected).
Synthetic dataset build tests (all 1M vector datasets)
Also tested real-world BIGANN 10M (uint8 128D) and GIST (fp32 960D) datasets: